from __future__ import annotations from datetime import datetime, timedelta, timezone from decimal import Decimal from fastapi import HTTPException from sqlalchemy import select, text from sqlalchemy.ext.asyncio import AsyncSession from app.enums.credit_balance import CreditBalanceSourceType, CreditLevel, CreditScope from app.enums.credit_product import CreditProductType from app.enums.credit_record import CreditRecordBillingScene, CreditRecordType from app.enums.credit_subscription import CreditSubscriptionPeriodStatus, CreditSubscriptionStatus from app.enums.team import TeamStatus from app.models.credit.balance import UserCreditBalance from app.models.credit.product import CreditProduct from app.models.credit.subscription import UserCreditSubscription from app.models.credit.subscription_period import UserCreditSubscriptionPeriod from app.models.payment_order import PaymentOrder from app.models.team import Team from app.models.user import User from app.services.credit.ledger_service import grant_credits, revoke_balances from app.services.credit.locking import ( acquire_subscription_credit_lock, acquire_team_business_lock, acquire_user_credit_lock, ) from app.services.credit.product_service import grant_count_for_cycle from app.services.credit.time_policy import add_natural_months, natural_month_period from app.services.credit.utils import to_credit_decimal, utc_now from app.services.credit_record_meta_service import CreditRecordMeta from app.services.operation_log_service import log_operation_event from app.services.team_service import create_team_for_subscription from app.utils.id_gen import generate_id _BUSINESS_TZ = timezone(timedelta(hours=8)) async def _generate_subscription_no( db: AsyncSession, *, product_type: str, created_at: datetime, ) -> str: sequence = int((await db.execute(text("SELECT nextval('credit_subscription_no_seq')"))).scalar_one()) prefix = "TS" if product_type == CreditProductType.TEAM_SUBSCRIPTION.value else "PS" normalized_at = created_at if created_at.tzinfo is not None else created_at.replace(tzinfo=timezone.utc) local_date = normalized_at.astimezone(_BUSINESS_TZ).strftime("%Y%m%d") return f"{prefix}{local_date}{sequence:06d}" def _allocate_amount_by_period(total: Decimal, count: int) -> list[Decimal]: value = to_credit_decimal(total) cents = int((value * 100).to_integral_value()) base, remainder = divmod(cents, max(1, count)) values = [Decimal(base) / 100 for _ in range(max(1, count))] values[-1] += Decimal(remainder) / 100 return [to_credit_decimal(item) for item in values] def _product_snapshot(product: CreditProduct) -> dict: return { "id": product.id, "product_code": product.product_code, "product_type": product.product_type, "name": product.name, "tier_code": product.tier_code, "tier_rank": product.tier_rank, "billing_cycle": product.billing_cycle, "monthly_grant_credits": float(product.monthly_grant_credits or 0), "first_purchase_price": float(product.first_purchase_price or 0), "regular_price": float(product.regular_price or 0), "activity_price": float(product.activity_price) if product.activity_price is not None else None, "grant_credits": float(product.grant_credits or 0), "validity_months": int(product.validity_months or 1) if product.is_credit_addon else None, "credit_level": product.credit_level, "features": product.features_json or [], } def _resolve_order_product_snapshot(order: PaymentOrder, product: CreditProduct | None) -> dict: snapshot = dict(order.product_snapshot_json or {}) if not snapshot and product is not None: snapshot = _product_snapshot(product) if not snapshot: raise ValueError("订单缺少商品快照,无法履约") return snapshot def _snapshot_decimal(snapshot: dict, key: str, default: Decimal = Decimal("0.00")) -> Decimal: if snapshot.get(key) is None: return default return to_credit_decimal(snapshot.get(key)) def _snapshot_validity_months(snapshot: dict) -> int: months = int(snapshot.get("validity_months") or 1) if not 1 <= months <= 36: raise ValueError("积分增值包有效期必须为1到36个月") return months def _actual_unit_price(order: PaymentOrder) -> Decimal: if order.actual_unit_price_snapshot is not None: return Decimal(str(order.actual_unit_price_snapshot)) quantity = max(1, int(order.quantity or 1)) return (Decimal(str(order.amount)) / Decimal(quantity)).quantize(Decimal("0.000001")) async def _load_product_for_order(db: AsyncSession, product_id: str | None) -> CreditProduct | None: if not product_id: return None result = await db.execute(select(CreditProduct).where(CreditProduct.id == product_id).limit(1)) return result.scalar_one_or_none() async def _resolve_team_for_order( db: AsyncSession, *, order: PaymentOrder, user: User, fulfilled_at: datetime, ) -> Team: if order.team_id_snapshot: await acquire_team_business_lock(db, order.team_id_snapshot) result = await db.execute( select(Team) .where(Team.id == order.team_id_snapshot, Team.deleted_at.is_(None)) .limit(1) .with_for_update() ) team = result.scalar_one_or_none() if not team: raise HTTPException(status_code=409, detail="订单绑定的团队不存在") if team.manager_id != user.id: raise HTTPException(status_code=409, detail="当前用户已不是该团队队长,不能履约团队订阅订单") if team.status != TeamStatus.ACTIVE.value: raise HTTPException(status_code=409, detail="团队已禁用,不能履约新的团队订阅订单") return team if user.team_id: # 创建订单时没有团队快照,说明这是“无团队首购”。订单存续期间加入团队本应被业务层拦截; # 若出现旁路变更,必须拒绝履约,不能把旧订单动态绑定到后来加入的团队。 raise HTTPException(status_code=409, detail="创建订单后团队关系已发生变化,不能履约该团队订阅订单") team = await create_team_for_subscription(db, manager_user=user, started_at=fulfilled_at) order.team_id_snapshot = team.id return team async def _create_subscription( db: AsyncSession, *, order: PaymentOrder, product: CreditProduct | None, paid_at: datetime, team: Team | None = None, ) -> UserCreditSubscription: snapshot = _resolve_order_product_snapshot(order, product) product_type = str(order.product_type or snapshot.get("product_type") or "") if product_type not in { CreditProductType.SUBSCRIPTION.value, CreditProductType.TEAM_SUBSCRIPTION.value, }: raise ValueError("订单不是订阅套餐") billing_cycle = str(snapshot.get("billing_cycle") or "") grant_count = grant_count_for_cycle(billing_cycle) monthly_grant = _snapshot_decimal(snapshot, "monthly_grant_credits") if monthly_grant <= 0: raise ValueError("订阅套餐月度积分快照无效") quantity = int(order.quantity or 1) if product_type == CreditProductType.SUBSCRIPTION.value: quantity = 1 elif not 2 <= quantity <= 1000: raise ValueError("团队订阅数量必须在2到1000之间") monthly_total = to_credit_decimal(monthly_grant * quantity) expires_at = add_natural_months(paid_at, grant_count) allocated_paid = _allocate_amount_by_period(to_credit_decimal(order.amount), grant_count) subscription_no = await _generate_subscription_no( db, product_type=product_type, created_at=paid_at ) subscription = UserCreditSubscription( id=generate_id(), subscription_no=subscription_no, user_id=order.user_id, team_id=team.id if team else None, team_manager_id_snapshot=order.user_id if team else None, product_id=order.product_id, payment_order_id=order.id, status=CreditSubscriptionStatus.ACTIVE.value, purchase_scene=order.purchase_scene or order.price_type or "regular", product_type_snapshot=product_type, product_name_snapshot=str(order.product_name_snapshot or snapshot.get("name") or "订阅套餐"), tier_code=str(snapshot.get("tier_code") or "starter"), tier_rank=int(snapshot.get("tier_rank") or 0), billing_cycle=billing_cycle, anchor_at=paid_at, start_at=paid_at, expires_at=expires_at, next_grant_at=paid_at, monthly_grant_credits_snapshot=monthly_grant, monthly_total_credits_snapshot=monthly_total, quantity_snapshot=quantity, grant_count=grant_count, granted_count=0, first_purchase_price_snapshot=_snapshot_decimal(snapshot, "first_purchase_price"), regular_price_snapshot=_snapshot_decimal(snapshot, "regular_price"), activity_price_snapshot=( _snapshot_decimal(snapshot, "activity_price") if snapshot.get("activity_price") is not None else None ), actual_unit_price_snapshot=_actual_unit_price(order), paid_amount_snapshot=to_credit_decimal(order.amount), product_snapshot_json=snapshot, ) db.add(subscription) await db.flush() for sequence in range(grant_count): valid_from, period_end = natural_month_period(paid_at, sequence) db.add( UserCreditSubscriptionPeriod( id=generate_id(), subscription_id=subscription.id, sequence=sequence, scheduled_at=valid_from, valid_from=valid_from, expires_at=min(period_end, expires_at), grant_credits=monthly_total, allocated_paid_amount=allocated_paid[sequence], status=CreditSubscriptionPeriodStatus.SCHEDULED.value, ) ) await db.flush() return subscription async def grant_subscription_period( db: AsyncSession, *, subscription: UserCreditSubscription, period: UserCreditSubscriptionPeriod, request_time: datetime | None = None, ) -> UserCreditBalance | None: checked_at = request_time or utc_now() await acquire_subscription_credit_lock(db, subscription.id) if period.status == CreditSubscriptionPeriodStatus.GRANTED.value and period.issued_balance_id: result = await db.execute(select(UserCreditBalance).where(UserCreditBalance.id == period.issued_balance_id)) return result.scalar_one_or_none() if subscription.status != CreditSubscriptionStatus.ACTIVE.value or checked_at >= subscription.expires_at: return None if period.valid_from > checked_at: return None is_team = subscription.product_type_snapshot == CreditProductType.TEAM_SUBSCRIPTION.value scope = CreditScope.TEAM.value if is_team else CreditScope.PERSONAL.value credit_level = str((subscription.product_snapshot_json or {}).get("credit_level") or CreditLevel.GENERAL.value) result = await grant_credits( db, user_id=subscription.user_id, amount=period.grant_credits, description=( f"团队订阅《{subscription.product_name_snapshot}》第{period.sequence + 1}期积分发放" if is_team else f"个人订阅《{subscription.product_name_snapshot}》第{period.sequence + 1}期积分发放" ), source_type=CreditBalanceSourceType.SUBSCRIPTION_GRANT.value, valid_from=period.valid_from, expires_at=period.expires_at, credit_level=credit_level, source_id=subscription.id, product_id=subscription.product_id, payment_order_id=subscription.payment_order_id, subscription_id=subscription.id, subscription_period_id=period.id, related_id=subscription.id, record_type=CreditRecordType.RECHARGE.value, biz_key=f"subscription:{subscription.id}:period:{period.id}:grant", record_meta=CreditRecordMeta(billing_scene=CreditRecordBillingScene.SUBSCRIPTION_GRANT.value), metadata_json={ "subscription_id": subscription.id, "subscription_period_id": period.id, "sequence": period.sequence, "product_type": subscription.product_type_snapshot, "quantity": subscription.quantity_snapshot, }, request_time=checked_at, credit_scope=scope, team_id=subscription.team_id, team_manager_id_snapshot=subscription.team_manager_id_snapshot, ) if not result.record: return None balance_result = await db.execute( select(UserCreditBalance).where(UserCreditBalance.grant_record_id == result.record.id).limit(1) ) balance = balance_result.scalar_one_or_none() if not balance: raise RuntimeError("订阅积分发放后未找到对应余额批次") period.issued_balance_id = balance.id period.issued_at = checked_at period.status = CreditSubscriptionPeriodStatus.GRANTED.value subscription.granted_count = max(int(subscription.granted_count or 0), period.sequence + 1) subscription.next_grant_at = ( add_natural_months(subscription.anchor_at, period.sequence + 1) if period.sequence + 1 < subscription.grant_count else None ) await db.flush() log_operation_event( domain="credit", module="subscription", event_type="SUBSCRIPTION_PERIOD_GRANTED", user_id=subscription.user_id, message="订阅周期积分发放成功", detail={ "subscription_id": subscription.id, "subscription_no": subscription.subscription_no, "period_id": period.id, "team_id": subscription.team_id, "quantity": subscription.quantity_snapshot, "amount": float(period.grant_credits), }, ) return balance async def fulfill_payment_product( db: AsyncSession, *, order: PaymentOrder, fulfilled_at: datetime | None = None, ) -> dict: checked_at = fulfilled_at or order.paid_at or utc_now() if order.status == "refunded": raise HTTPException(status_code=409, detail="订单已退款,不能再履约") if order.fulfillment_status == "fulfilled": return {"fulfilled": True, "subscription_id": order.subscription_id, "created": False} await acquire_user_credit_lock(db, order.user_id) product = await _load_product_for_order(db, order.product_id) snapshot = _resolve_order_product_snapshot(order, product) product_type = str(order.product_type or snapshot.get("product_type") or "") # 固定锁序:user -> team -> subscription/row。已绑定团队的订单先拿 Team advisory, # 避免履约与 Seat/消费并发时出现 User row -> Team 的反向锁序。 if product_type == CreditProductType.TEAM_SUBSCRIPTION.value and order.team_id_snapshot: await acquire_team_business_lock(db, order.team_id_snapshot) user_result = await db.execute(select(User).where(User.id == order.user_id).limit(1).with_for_update()) user = user_result.scalar_one_or_none() if not user: raise ValueError("订单用户不存在") if product_type == CreditProductType.CREDIT_ADDON.value: months = _snapshot_validity_months(snapshot) grant_amount = _snapshot_decimal(snapshot, "grant_credits") if grant_amount <= 0: raise ValueError("积分增值包积分快照无效") await grant_credits( db, user_id=order.user_id, amount=grant_amount, description=f"购买积分增值包《{order.product_name_snapshot or snapshot.get('name') or '积分增值包'}》", source_type=CreditBalanceSourceType.CREDIT_ADDON.value, valid_from=checked_at, expires_at=add_natural_months(checked_at, months), credit_level=str(snapshot.get("credit_level") or CreditLevel.GENERAL.value), source_id=order.id, product_id=order.product_id, payment_order_id=order.id, related_id=order.id, record_type=CreditRecordType.RECHARGE.value, biz_key=f"payment:{order.id}:credit-addon", record_meta=CreditRecordMeta(billing_scene=CreditRecordBillingScene.CREDIT_ADDON_PURCHASE.value), metadata_json={"payment_order_id": order.id, "product_snapshot": snapshot}, request_time=checked_at, ) order.fulfillment_status = "fulfilled" order.fulfilled_at = checked_at order.subscription_id = None await db.flush() return {"fulfilled": True, "subscription_id": None, "created": True} if product_type not in { CreditProductType.SUBSCRIPTION.value, CreditProductType.TEAM_SUBSCRIPTION.value, }: raise ValueError("订单商品类型不支持履约") existing_result = await db.execute( select(UserCreditSubscription).where(UserCreditSubscription.payment_order_id == order.id).limit(1) ) existing = existing_result.scalar_one_or_none() if existing: order.subscription_id = existing.id order.fulfillment_status = "fulfilled" order.fulfilled_at = order.fulfilled_at or checked_at return {"fulfilled": True, "subscription_id": existing.id, "created": False} team: Team | None = None if product_type == CreditProductType.TEAM_SUBSCRIPTION.value: team = await _resolve_team_for_order(db, order=order, user=user, fulfilled_at=checked_at) subscription = await _create_subscription( db, order=order, product=product, paid_at=checked_at, team=team ) first_period_result = await db.execute( select(UserCreditSubscriptionPeriod) .where(UserCreditSubscriptionPeriod.subscription_id == subscription.id) .order_by(UserCreditSubscriptionPeriod.sequence.asc()) .limit(1) ) first_period = first_period_result.scalar_one() await grant_subscription_period(db, subscription=subscription, period=first_period, request_time=checked_at) if product_type == CreditProductType.SUBSCRIPTION.value and user.first_membership_paid_at is None: user.first_membership_paid_at = checked_at if team is not None and team.first_subscription_paid_at is None: team.first_subscription_paid_at = checked_at order.subscription_id = subscription.id order.fulfillment_status = "fulfilled" order.fulfilled_at = checked_at await db.flush() log_operation_event( domain="payment", module="subscription", event_type="SUBSCRIPTION_FULFILLED", user_id=order.user_id, message="订阅订单履约成功", detail={ "order_no": order.order_no, "subscription_id": subscription.id, "subscription_no": subscription.subscription_no, "product_type": product_type, "team_id": subscription.team_id, "quantity": subscription.quantity_snapshot, "paid_amount": float(subscription.paid_amount_snapshot), }, ) return {"fulfilled": True, "subscription_id": subscription.id, "created": True} async def list_due_subscription_period_candidates( db: AsyncSession, *, request_time: datetime, batch_size: int = 100, ) -> list[tuple[str, str, str, str | None]]: result = await db.execute( select( UserCreditSubscriptionPeriod.id, UserCreditSubscription.id.label("subscription_id"), UserCreditSubscription.user_id, UserCreditSubscription.team_id, ) .join(UserCreditSubscription, UserCreditSubscription.id == UserCreditSubscriptionPeriod.subscription_id) .where( UserCreditSubscription.status == CreditSubscriptionStatus.ACTIVE.value, UserCreditSubscription.expires_at > request_time, UserCreditSubscriptionPeriod.status == CreditSubscriptionPeriodStatus.SCHEDULED.value, UserCreditSubscriptionPeriod.scheduled_at <= request_time, ) .order_by(UserCreditSubscriptionPeriod.scheduled_at.asc(), UserCreditSubscriptionPeriod.id.asc()) .limit(batch_size) ) return [(str(r.id), str(r.subscription_id), str(r.user_id), str(r.team_id) if r.team_id else None) for r in result.all()] async def grant_due_subscription_period_by_id( db: AsyncSession, *, period_id: str, subscription_id: str, user_id: str, team_id: str | None, request_time: datetime, ) -> bool: await acquire_user_credit_lock(db, user_id) if team_id: await acquire_team_business_lock(db, team_id) await acquire_subscription_credit_lock(db, subscription_id) sub_result = await db.execute( select(UserCreditSubscription).where(UserCreditSubscription.id == subscription_id).limit(1).with_for_update() ) subscription = sub_result.scalar_one_or_none() period_result = await db.execute( select(UserCreditSubscriptionPeriod).where(UserCreditSubscriptionPeriod.id == period_id).limit(1).with_for_update() ) period = period_result.scalar_one_or_none() if not subscription or not period: return False balance = await grant_subscription_period(db, subscription=subscription, period=period, request_time=request_time) return balance is not None async def grant_due_subscription_periods( db: AsyncSession, *, request_time: datetime | None = None, batch_size: int = 100, ) -> int: checked_at = request_time or utc_now() candidates = await list_due_subscription_period_candidates(db, request_time=checked_at, batch_size=batch_size) count = 0 for period_id, subscription_id, user_id, team_id in candidates: if await grant_due_subscription_period_by_id( db, period_id=period_id, subscription_id=subscription_id, user_id=user_id, team_id=team_id, request_time=checked_at, ): count += 1 return count async def list_due_subscription_ids( db: AsyncSession, *, request_time: datetime, batch_size: int = 500, ) -> list[str]: result = await db.execute( select(UserCreditSubscription.id) .where( UserCreditSubscription.status == CreditSubscriptionStatus.ACTIVE.value, UserCreditSubscription.expires_at <= request_time, ) .order_by(UserCreditSubscription.expires_at.asc(), UserCreditSubscription.id.asc()) .limit(batch_size) ) return [str(item) for item in result.scalars().all()] async def expire_subscription_by_id( db: AsyncSession, *, subscription_id: str, request_time: datetime, ) -> bool: await acquire_subscription_credit_lock(db, subscription_id) result = await db.execute( select(UserCreditSubscription).where(UserCreditSubscription.id == subscription_id).limit(1).with_for_update() ) subscription = result.scalar_one_or_none() if not subscription or subscription.status != CreditSubscriptionStatus.ACTIVE.value: return False if subscription.expires_at > request_time: return False subscription.status = CreditSubscriptionStatus.EXPIRED.value subscription.next_grant_at = None periods_result = await db.execute( select(UserCreditSubscriptionPeriod).where( UserCreditSubscriptionPeriod.subscription_id == subscription.id, UserCreditSubscriptionPeriod.status == CreditSubscriptionPeriodStatus.SCHEDULED.value, ) ) for period in periods_result.scalars().all(): period.status = CreditSubscriptionPeriodStatus.EXPIRED.value period.cancelled_at = request_time await db.flush() return True async def expire_due_subscriptions( db: AsyncSession, *, request_time: datetime | None = None, batch_size: int = 500, ) -> int: checked_at = request_time or utc_now() ids = await list_due_subscription_ids(db, request_time=checked_at, batch_size=batch_size) count = 0 for subscription_id in ids: if await expire_subscription_by_id(db, subscription_id=subscription_id, request_time=checked_at): count += 1 return count async def revoke_payment_order_credits( db: AsyncSession, *, order: PaymentOrder, reason: str, request_time: datetime | None = None, ) -> None: """保留历史兼容函数;本版本主动订单退款入口关闭,不应由新退款流程调用。""" checked_at = request_time or utc_now() balances_result = await db.execute( select(UserCreditBalance).where( UserCreditBalance.payment_order_id == order.id, UserCreditBalance.revoked_at.is_(None), UserCreditBalance.unspent_amount > 0, ) ) balances = list(balances_result.scalars().all()) if balances: await revoke_balances( db, balances=balances, description=reason, related_id=order.id, biz_key=f"payment:{order.id}:legacy-revoke", request_time=checked_at, )