199 lines
6.5 KiB
Python
199 lines
6.5 KiB
Python
"""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("分身尚未关联有效用户,暂时无法使用积分")
|
|
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("积分余额不足,请充值后继续")
|
|
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("积分账户不存在")
|
|
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)
|