236 lines
9.8 KiB
Python
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)
|