Files
video-gen/video-gen-api/app/services/credit/subscription_service.py
T
2026-08-14 15:23:46 +08:00

608 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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,
)