"""User-scoped token accounting for every avatar model request.""" import math from dataclasses import dataclass from datetime import datetime, timedelta from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from models import Avatar, TokenAccount, TokenUsage, User DEFAULT_TOKEN_GRANT = 1_000_000 class InsufficientTokensError(RuntimeError): pass @dataclass(frozen=True) class TokenReservation: usage_id: str user_id: str reserved_tokens: int def get_or_create_account(db: Session, user_id: str) -> TokenAccount: account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first() if account: return account account = TokenAccount( user_id=user_id, balance=DEFAULT_TOKEN_GRANT, total_granted=DEFAULT_TOKEN_GRANT, total_consumed=0, ) db.add(account) try: db.commit() except IntegrityError: # A concurrent first request may have created the same user account. db.rollback() account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first() if account is None: raise db.refresh(account) return account def avatar_owner_user(db: Session, avatar: Avatar) -> User | None: owner_id = (avatar.owner_id or "").strip() if not owner_id: return None return db.query(User).filter(User.huihui_user_id == owner_id).first() def estimate_request_tokens(messages: list[dict], max_output_tokens: int) -> int: # UTF-8 bytes / 2 deliberately overestimates mixed Chinese/English prompts; # the unused reservation is returned after provider usage is received. content_bytes = sum( len(str(item.get("content", "")).encode("utf-8")) for item in messages ) prompt_reserve = max(1, math.ceil(content_bytes / 2) + len(messages) * 6) return prompt_reserve + max(1, int(max_output_tokens)) def estimate_fallback_usage(messages: list[dict], output: str) -> int: content_bytes = sum( len(str(item.get("content", "")).encode("utf-8")) for item in messages ) + len((output or "").encode("utf-8")) return max(1, math.ceil(content_bytes / 3) + len(messages) * 4) def reserve_avatar_tokens( db: Session, avatar: Avatar, source: str, model: str, messages: list[dict], max_output_tokens: int, ) -> TokenReservation: user = avatar_owner_user(db, avatar) if not user: raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用 Token") account = get_or_create_account(db, user.id) reserved = estimate_request_tokens(messages, max_output_tokens) updated = ( db.query(TokenAccount) .filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved) .update( {TokenAccount.balance: TokenAccount.balance - reserved}, synchronize_session=False, ) ) if updated != 1: db.rollback() raise InsufficientTokensError("Token 余额不足,请充值后继续") db.refresh(account) usage = TokenUsage( user_id=user.id, avatar_id=avatar.id, source=source, model=model, status="reserved", reserved_tokens=reserved, ) db.add(usage) db.flush() usage.balance_after = account.balance db.commit() return TokenReservation(usage.id, user.id, reserved) def settle_reservation( db: Session, reservation: TokenReservation, usage: dict | None, *, fallback_total: int, ) -> dict: record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() if not record or record.status != "reserved": return {} provider_usage = usage or {} prompt_tokens = max(0, int(provider_usage.get("prompt_tokens") or 0)) completion_tokens = max(0, int(provider_usage.get("completion_tokens") or 0)) provider_total = max( int(provider_usage.get("total_tokens") or 0), prompt_tokens + completion_tokens, ) total_tokens = max(1, provider_total or int(fallback_total or 0)) updated = ( db.query(TokenAccount) .filter(TokenAccount.user_id == reservation.user_id) .update( { TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens - total_tokens, TokenAccount.total_consumed: TokenAccount.total_consumed + total_tokens, }, synchronize_session=False, ) ) if updated != 1: raise RuntimeError("Token 账户不存在") db.expire_all() account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first() record.prompt_tokens = prompt_tokens record.completion_tokens = completion_tokens record.total_tokens = total_tokens record.balance_after = account.balance record.status = "completed" record.settled_at = datetime.utcnow() db.commit() return { "promptTokens": prompt_tokens, "completionTokens": completion_tokens, "totalTokens": total_tokens, "balance": account.balance, } def release_reservation(db: Session, reservation: TokenReservation, reason: str = "") -> None: record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() if not record or record.status != "reserved": return updated = ( db.query(TokenAccount) .filter(TokenAccount.user_id == reservation.user_id) .update( {TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens}, synchronize_session=False, ) ) if updated: db.expire_all() account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first() record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() record.balance_after = account.balance record.status = "failed" record.failure_reason = (reason or "model_request_failed")[:255] record.settled_at = datetime.utcnow() db.commit() def release_stale_reservations(db: Session, older_than_minutes: int = 10) -> int: cutoff = datetime.utcnow() - timedelta(minutes=older_than_minutes) stale = db.query(TokenUsage).filter( TokenUsage.status == "reserved", TokenUsage.created_at < cutoff, ).all() for record in stale: release_reservation( db, TokenReservation(record.id, record.user_id, int(record.reserved_tokens or 0)), "stale_reservation_recovered", ) return len(stale)