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

236 lines
9.8 KiB
Python

from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime
from decimal import Decimal
from typing import Iterable
from sqlalchemy import and_, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.credit_balance import CreditBalanceStatus, CreditScope
from app.enums.credit_subscription import CreditSubscriptionStatus
from app.enums.team import TeamStatus
from app.models.credit.balance import UserCreditBalance
from app.models.credit.subscription import UserCreditSubscription
from app.models.credit.subscription_period import UserCreditSubscriptionPeriod
from app.models.credit.team_seat import TeamSubscriptionSeat
from app.models.credit.team_seat_usage import TeamSubscriptionSeatUsage
from app.models.team import Team
from app.models.user import User
from app.services.credit.time_policy import last_usable_at
from app.services.credit.utils import to_credit_decimal, to_float, utc_now
@dataclass(slots=True, frozen=True)
class CreditBalanceSummary:
personal_credits: Decimal
team_available_credits: Decimal
team_frozen_credits: Decimal
available_credits: Decimal
next_expiring_credits: Decimal
next_expires_at: datetime | None
next_last_usable_at: datetime | None
def to_dict(self) -> dict:
return {
"personal_credits": to_float(self.personal_credits),
"team_available_credits": to_float(self.team_available_credits),
"team_frozen_credits": to_float(self.team_frozen_credits),
"available_credits": to_float(self.available_credits),
"credits": to_float(self.available_credits),
"next_expiring_credits": to_float(self.next_expiring_credits),
"next_expires_at": self.next_expires_at,
"next_last_usable_at": self.next_last_usable_at,
}
def _blank_summary() -> CreditBalanceSummary:
zero = Decimal("0.00")
return CreditBalanceSummary(zero, zero, zero, zero, zero, None, None)
async def get_user_credit_summary_map(
db: AsyncSession,
user_ids: Iterable[str],
*,
request_time: datetime | None = None,
) -> dict[str, CreditBalanceSummary]:
ids = list(dict.fromkeys(str(item) for item in user_ids if item))
if not ids:
return {}
checked_at = request_time or utc_now()
personal_map = {user_id: Decimal("0.00") for user_id in ids}
team_available_map = {user_id: Decimal("0.00") for user_id in ids}
team_frozen_map = {user_id: Decimal("0.00") for user_id in ids}
expiry_buckets: dict[str, dict[datetime, Decimal]] = {user_id: {} for user_id in ids}
personal_result = await db.execute(
select(UserCreditBalance.user_id, UserCreditBalance.expires_at, UserCreditBalance.unspent_amount)
.where(
UserCreditBalance.user_id.in_(ids),
UserCreditBalance.credit_scope == CreditScope.PERSONAL.value,
UserCreditBalance.valid_from <= checked_at,
UserCreditBalance.expires_at > checked_at,
UserCreditBalance.unspent_amount > 0,
UserCreditBalance.revoked_at.is_(None),
)
)
for row in personal_result.all():
amount = to_credit_decimal(row.unspent_amount)
user_id = str(row.user_id)
personal_map[user_id] += amount
expiry_buckets[user_id][row.expires_at] = expiry_buckets[user_id].get(row.expires_at, Decimal("0.00")) + amount
team_result = await db.execute(
select(
TeamSubscriptionSeat.user_id,
Team.status.label("team_status"),
UserCreditBalance.expires_at,
UserCreditBalance.unspent_amount,
TeamSubscriptionSeat.monthly_allocated_credits,
TeamSubscriptionSeatUsage.used_credits,
)
.join(UserCreditSubscription, UserCreditSubscription.id == TeamSubscriptionSeat.subscription_id)
.join(Team, Team.id == TeamSubscriptionSeat.team_id)
.join(
UserCreditSubscriptionPeriod,
UserCreditSubscriptionPeriod.subscription_id == UserCreditSubscription.id,
)
.join(UserCreditBalance, UserCreditBalance.id == UserCreditSubscriptionPeriod.issued_balance_id)
.outerjoin(
TeamSubscriptionSeatUsage,
(TeamSubscriptionSeatUsage.seat_id == TeamSubscriptionSeat.id)
& (TeamSubscriptionSeatUsage.subscription_period_id == UserCreditSubscriptionPeriod.id),
)
.where(
TeamSubscriptionSeat.user_id.in_(ids),
TeamSubscriptionSeat.deleted_at.is_(None),
TeamSubscriptionSeat.cancelled_at.is_(None),
Team.deleted_at.is_(None),
UserCreditSubscription.status == CreditSubscriptionStatus.ACTIVE.value,
UserCreditSubscription.start_at <= checked_at,
UserCreditSubscription.expires_at > checked_at,
UserCreditSubscriptionPeriod.valid_from <= checked_at,
UserCreditSubscriptionPeriod.expires_at > checked_at,
UserCreditBalance.credit_scope == CreditScope.TEAM.value,
UserCreditBalance.valid_from <= checked_at,
UserCreditBalance.expires_at > checked_at,
UserCreditBalance.unspent_amount > 0,
UserCreditBalance.revoked_at.is_(None),
)
)
for row in team_result.all():
allocated = to_credit_decimal(row.monthly_allocated_credits)
used = to_credit_decimal(row.used_credits or 0)
seat_remaining = max(Decimal("0.00"), allocated - used)
amount = min(to_credit_decimal(row.unspent_amount), seat_remaining)
if amount <= 0:
continue
user_id = str(row.user_id)
if row.team_status == TeamStatus.ACTIVE.value:
team_available_map[user_id] += amount
expiry_buckets[user_id][row.expires_at] = expiry_buckets[user_id].get(row.expires_at, Decimal("0.00")) + amount
else:
team_frozen_map[user_id] += amount
output: dict[str, CreditBalanceSummary] = {}
for user_id in ids:
personal = personal_map[user_id]
team_available = team_available_map[user_id]
frozen = team_frozen_map[user_id]
available = personal + team_available
bucket = expiry_buckets[user_id]
next_expires_at = min(bucket.keys()) if bucket else None
next_expiring = bucket.get(next_expires_at, Decimal("0.00")) if next_expires_at else Decimal("0.00")
output[user_id] = CreditBalanceSummary(
personal_credits=personal,
team_available_credits=team_available,
team_frozen_credits=frozen,
available_credits=available,
next_expiring_credits=next_expiring,
next_expires_at=next_expires_at,
next_last_usable_at=last_usable_at(next_expires_at) if next_expires_at else None,
)
return output
async def get_balance_summary(
db: AsyncSession,
user_id: str,
*,
request_time: datetime | None = None,
) -> CreditBalanceSummary:
return (await get_user_credit_summary_map(db, [user_id], request_time=request_time)).get(
user_id, _blank_summary()
)
async def get_available_credits(
db: AsyncSession,
user_id: str,
*,
request_time: datetime | None = None,
) -> Decimal:
return (await get_balance_summary(db, user_id, request_time=request_time)).available_credits
async def get_user_credit_map(
db: AsyncSession,
user_ids: Iterable[str],
*,
request_time: datetime | None = None,
) -> dict[str, float]:
ids = list(dict.fromkeys(str(item) for item in user_ids if item))
summaries = await get_user_credit_summary_map(db, ids, request_time=request_time)
return {user_id: to_float(summaries.get(user_id, _blank_summary()).available_credits) for user_id in ids}
def attach_credit_snapshot(user: object, credits: Decimal | float | int) -> object:
setattr(user, "credits", to_float(to_credit_decimal(credits)))
return user
def effective_balance_status(balance: UserCreditBalance, *, request_time: datetime | None = None) -> str:
checked_at = request_time or utc_now()
if balance.revoked_at is not None or to_credit_decimal(balance.revoked_amount) > 0:
return CreditBalanceStatus.REVOKED.value
if balance.expires_at <= checked_at:
return CreditBalanceStatus.EXPIRED.value
if balance.valid_from > checked_at:
return CreditBalanceStatus.SCHEDULED.value
if to_credit_decimal(balance.unspent_amount) <= 0:
return CreditBalanceStatus.CONSUMED.value
return CreditBalanceStatus.ACTIVE.value
def apply_balance_status_filter(stmt, status: str | None, *, request_time: datetime):
if not status:
return stmt
if status == CreditBalanceStatus.REVOKED.value:
return stmt.where(or_(UserCreditBalance.revoked_at.is_not(None), UserCreditBalance.revoked_amount > 0))
base_not_revoked = and_(UserCreditBalance.revoked_at.is_(None), UserCreditBalance.revoked_amount <= 0)
if status == CreditBalanceStatus.EXPIRED.value:
return stmt.where(base_not_revoked, UserCreditBalance.expires_at <= request_time)
if status == CreditBalanceStatus.SCHEDULED.value:
return stmt.where(
base_not_revoked,
UserCreditBalance.valid_from > request_time,
UserCreditBalance.expires_at > request_time,
)
if status == CreditBalanceStatus.CONSUMED.value:
return stmt.where(
base_not_revoked,
UserCreditBalance.valid_from <= request_time,
UserCreditBalance.expires_at > request_time,
UserCreditBalance.unspent_amount <= 0,
)
if status == CreditBalanceStatus.ACTIVE.value:
return stmt.where(
base_not_revoked,
UserCreditBalance.valid_from <= request_time,
UserCreditBalance.expires_at > request_time,
UserCreditBalance.unspent_amount > 0,
)
return stmt.where(UserCreditBalance.status == status)