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

113 lines
4.1 KiB
Python

from __future__ import annotations
from datetime import datetime
from sqlalchemy import case, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.credit_product import CreditProductType, SubscriptionBillingCycle
from app.enums.credit_subscription import CreditSubscriptionStatus
from app.enums.team import TeamStatus
from app.models.credit.product import CreditProduct
from app.models.credit.subscription import UserCreditSubscription
from app.models.credit.team_seat import TeamSubscriptionSeat
from app.models.team import Team
from app.services.credit.utils import utc_now
_BILLING_CYCLE_RANK = case(
(CreditProduct.billing_cycle == SubscriptionBillingCycle.YEARLY.value, 3),
(CreditProduct.billing_cycle == SubscriptionBillingCycle.QUARTERLY.value, 2),
(CreditProduct.billing_cycle == SubscriptionBillingCycle.MONTHLY.value, 1),
else_=0,
)
def _entitlement_payload(subscription: UserCreditSubscription, product: CreditProduct) -> dict:
return {
"subscription_id": subscription.id,
"product_id": product.id,
"product_name": product.name,
"product_type": subscription.product_type_snapshot,
"billing_cycle": product.billing_cycle,
"tier_code": product.tier_code,
"tier_rank": product.tier_rank,
"sort_order": product.sort_order,
"start_at": subscription.start_at,
"expires_at": subscription.expires_at,
"team_id": subscription.team_id,
}
async def get_personal_entitlement(
db: AsyncSession,
*,
user_id: str,
request_time: datetime | None = None,
) -> dict | None:
checked_at = request_time or utc_now()
result = await db.execute(
select(UserCreditSubscription, CreditProduct)
.join(CreditProduct, CreditProduct.id == UserCreditSubscription.product_id)
.where(
UserCreditSubscription.user_id == user_id,
UserCreditSubscription.product_type_snapshot == CreditProductType.SUBSCRIPTION.value,
UserCreditSubscription.status == CreditSubscriptionStatus.ACTIVE.value,
UserCreditSubscription.start_at <= checked_at,
UserCreditSubscription.expires_at > checked_at,
)
.order_by(
_BILLING_CYCLE_RANK.desc(),
CreditProduct.tier_rank.desc(),
CreditProduct.sort_order.asc(),
CreditProduct.created_at.desc(),
CreditProduct.id.desc(),
)
.limit(1)
)
row = result.first()
return _entitlement_payload(row[0], row[1]) if row else None
async def get_team_entitlement(
db: AsyncSession,
*,
user_id: str,
request_time: datetime | None = None,
) -> dict | None:
checked_at = request_time or utc_now()
result = await db.execute(
select(UserCreditSubscription, CreditProduct, TeamSubscriptionSeat)
.join(CreditProduct, CreditProduct.id == UserCreditSubscription.product_id)
.join(
TeamSubscriptionSeat,
TeamSubscriptionSeat.subscription_id == UserCreditSubscription.id,
)
.join(Team, Team.id == UserCreditSubscription.team_id)
.where(
TeamSubscriptionSeat.user_id == user_id,
TeamSubscriptionSeat.deleted_at.is_(None),
TeamSubscriptionSeat.cancelled_at.is_(None),
UserCreditSubscription.product_type_snapshot == CreditProductType.TEAM_SUBSCRIPTION.value,
UserCreditSubscription.status == CreditSubscriptionStatus.ACTIVE.value,
UserCreditSubscription.start_at <= checked_at,
UserCreditSubscription.expires_at > checked_at,
Team.deleted_at.is_(None),
Team.status == TeamStatus.ACTIVE.value,
)
.order_by(
_BILLING_CYCLE_RANK.desc(),
CreditProduct.tier_rank.desc(),
CreditProduct.sort_order.asc(),
CreditProduct.created_at.desc(),
CreditProduct.id.desc(),
)
.limit(1)
)
row = result.first()
if not row:
return None
payload = _entitlement_payload(row[0], row[1])
payload["seat_id"] = row[2].id
return payload