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)