feat(avatar): add user token accounting

This commit is contained in:
stefanfeng
2026-08-25 13:24:02 +08:00
parent 7a0199e685
commit 699bbbde57
15 changed files with 798 additions and 52 deletions
+14
View File
@@ -40,8 +40,14 @@ def init_db():
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 30"),
("avatars", "share_token", "VARCHAR DEFAULT NULL"),
("token_account", "user_id", "VARCHAR DEFAULT ''"),
("token_account", "total_granted", "BIGINT DEFAULT 0"),
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
("token_account", "created_at", "TIMESTAMP"),
("token_account", "updated_at", "TIMESTAMP"),
)
_normalize_optional_unique_values()
_create_token_indexes()
def _try_add_columns(*cols):
@@ -58,3 +64,11 @@ def _try_add_columns(*cols):
def _normalize_optional_unique_values():
with engine.begin() as conn:
conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''")
def _create_token_indexes():
with engine.begin() as conn:
conn.exec_driver_sql(
"CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id "
"ON token_account(user_id) WHERE user_id <> ''"
)
+25 -11
View File
@@ -8,7 +8,7 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.interval import IntervalTrigger
from database import init_db, SessionLocal
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan, User
from fastapi.staticfiles import StaticFiles
import routers.avatars
import routers.tokens
@@ -19,6 +19,7 @@ import routers.huihui_auth
import routers.chat
import routers.takeover
from responses import ok
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
logger = logging.getLogger(__name__)
@@ -56,17 +57,29 @@ def health():
def seed():
db = SessionLocal()
try:
if db.query(TokenAccount).first() is None:
db.add(TokenAccount(balance=1250))
plan_specs = [
{"id": "1", "name": "基础套餐", "amount": 2_000_000, "price": 10, "badge": "", "desc": "2M Token"},
{"id": "2", "name": "标准套餐", "amount": 20_000_000, "price": 100, "badge": "常用", "desc": "20M Token"},
{"id": "3", "name": "专业套餐", "amount": 250_000_000, "price": 1000, "badge": "加赠25%", "desc": "250M Token"},
{"id": "4", "name": "企业套餐", "amount": 2_500_000_000, "price": 10000, "badge": "企业推荐", "desc": "2500M Token"},
]
for spec in plan_specs:
plan = db.query(TokenPlan).filter(TokenPlan.id == spec["id"]).first()
if plan is None:
db.add(TokenPlan(**spec))
else:
for key, value in spec.items():
setattr(plan, key, value)
if db.query(TokenPlan).count() == 0:
plans = [
TokenPlan(id="1", name="新手体验", amount=1000, price=9.9, desc="新手体验"),
TokenPlan(id="2", name="热门套餐", amount=5000, price=39.9, badge="热门"),
TokenPlan(id="3", name="超值套餐", amount=12000, price=89.9, badge="超值"),
TokenPlan(id="4", name="企业推荐", amount=30000, price=199, badge="企业推荐", desc="适合高频使用"),
]
db.add_all(plans)
for user in db.query(User).all():
account = db.query(TokenAccount).filter(TokenAccount.user_id == user.id).first()
if account is None:
db.add(TokenAccount(
user_id=user.id,
balance=DEFAULT_TOKEN_GRANT,
total_granted=DEFAULT_TOKEN_GRANT,
total_consumed=0,
))
if db.query(Avatar).count() == 0:
avatar = Avatar(
@@ -105,6 +118,7 @@ def seed():
db.add_all(orgs)
db.commit()
release_stale_reservations(db)
finally:
db.close()
+31 -2
View File
@@ -1,6 +1,7 @@
import uuid
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
@@ -259,14 +260,42 @@ class KnowledgeChunk(Base):
class TokenAccount(Base):
__tablename__ = "token_account"
id = Column(Integer, primary_key=True)
balance = Column(Integer, default=1250)
user_id = Column(String, nullable=False, default="", index=True)
balance = Column(BigInteger, default=1_000_000)
total_granted = Column(BigInteger, default=1_000_000)
total_consumed = Column(BigInteger, default=0)
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
class TokenUsage(Base):
__tablename__ = "token_usage"
__table_args__ = (
Index("ix_token_usage_user_created", "user_id", "created_at"),
Index("ix_token_usage_avatar_created", "avatar_id", "created_at"),
)
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
user_id = Column(String, nullable=False, index=True)
avatar_id = Column(String, nullable=False, default="", index=True)
source = Column(String, nullable=False, default="chat")
model = Column(String, default="")
status = Column(String, nullable=False, default="reserved")
reserved_tokens = Column(BigInteger, default=0)
prompt_tokens = Column(BigInteger, default=0)
completion_tokens = Column(BigInteger, default=0)
total_tokens = Column(BigInteger, default=0)
balance_after = Column(BigInteger, default=0)
failure_reason = Column(String, default="")
created_at = Column(DateTime, server_default=func.now())
settled_at = Column(DateTime)
class TokenPlan(Base):
__tablename__ = "token_plans"
id = Column(String, primary_key=True)
name = Column(String, default="")
amount = Column(Integer, default=0)
amount = Column(BigInteger, default=0)
price = Column(Float, default=0)
badge = Column(String, default="")
desc = Column(String, default="")
+131 -15
View File
@@ -16,12 +16,20 @@ import embeddings
from database import get_db
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User
from responses import ok, fail
from services.token_billing import (
InsufficientTokensError,
estimate_fallback_usage,
release_reservation,
reserve_avatar_tokens,
settle_reservation,
)
router = APIRouter(tags=["数字分身聊天"])
CHAT_API_URL = os.getenv("CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1")
CHAT_API_KEY = os.getenv("CHAT_API_KEY", "")
CHAT_MODEL = os.getenv("CHAT_MODEL", "qwen-plus")
CHAT_MAX_OUTPUT_TOKENS = max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024")))
MAX_MESSAGE_LENGTH = 4000
MAX_HISTORY_MESSAGES = 10
QA_LEXICAL_THRESHOLD = 0.72
@@ -279,7 +287,7 @@ def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5
return results
def _call_qwen(messages: list[dict], temperature: float) -> str:
def _call_qwen(messages: list[dict], temperature: float) -> dict:
if not CHAT_API_KEY:
raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY")
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
@@ -287,6 +295,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str:
"model": CHAT_MODEL,
"messages": messages,
"temperature": temperature,
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
}
try:
response = httpx.post(
@@ -302,7 +311,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str:
raise RuntimeError("Qwen 模型服务暂时不可用") from exc
if not isinstance(answer, str) or not answer.strip():
raise RuntimeError("Qwen 模型没有返回有效回答")
return answer.strip()
return {"answer": answer.strip(), "usage": data.get("usage") or {}}
def _iter_qwen_stream(messages: list[dict], temperature: float):
@@ -310,7 +319,14 @@ def _iter_qwen_stream(messages: list[dict], temperature: float):
if not CHAT_API_KEY:
raise RuntimeError("模型服务未配置")
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
payload = {"model": CHAT_MODEL, "messages": messages, "temperature": temperature, "stream": True}
payload = {
"model": CHAT_MODEL,
"messages": messages,
"temperature": temperature,
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
"stream": True,
"stream_options": {"include_usage": True},
}
try:
with httpx.stream("POST", url, headers={"Authorization": f"Bearer {CHAT_API_KEY}"}, json=payload, timeout=45) as response:
response.raise_for_status()
@@ -322,11 +338,15 @@ def _iter_qwen_stream(messages: list[dict], temperature: float):
if data == "[DONE]":
return
try:
delta = json.loads(data).get("choices", [{}])[0].get("delta", {}).get("content")
parsed = json.loads(data)
except (ValueError, IndexError, AttributeError):
continue
if parsed.get("usage"):
yield {"usage": parsed["usage"]}
choices = parsed.get("choices") or []
delta = choices[0].get("delta", {}).get("content") if choices else None
if delta:
yield delta
yield {"content": delta}
except httpx.HTTPError as exc:
raise RuntimeError("模型服务暂时不可用") from exc
@@ -350,6 +370,7 @@ def _resolve_reply(
qa_pairs: list[Any] | None = None,
search_fn: Callable[..., list[dict]] | None = None,
model_client: Callable[..., str] | None = None,
usage_source: str = "chat",
) -> dict:
if qa_pairs is None:
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
@@ -362,16 +383,49 @@ def _resolve_reply(
messages = _build_prompt(avatar, history, question, hits)
config = _config(avatar)
temperature = min(0.45 if hits else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
model_client = model_client or _call_qwen
answer = model_client(messages=messages, temperature=temperature)
return {
token_usage = None
if model_client is not None:
answer = model_client(messages=messages, temperature=temperature)
else:
reservation = reserve_avatar_tokens(
db,
avatar,
usage_source,
CHAT_MODEL,
messages,
CHAT_MAX_OUTPUT_TOKENS,
)
try:
model_result = _call_qwen(messages=messages, temperature=temperature)
answer = model_result["answer"]
token_usage = settle_reservation(
db,
reservation,
model_result.get("usage"),
fallback_total=estimate_fallback_usage(messages, answer),
)
except Exception as exc:
release_reservation(db, reservation, str(exc))
raise
result = {
"answer": answer,
"source": "knowledge" if hits else "qwen",
"references": hits,
}
if token_usage:
result["tokenUsage"] = token_usage
return result
def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any], *, public: bool = False):
def _stream_reply(
db: Session,
avatar: Avatar,
question: str,
history: list[Any],
*,
public: bool = False,
usage_source: str = "chat_stream",
):
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
matched = _match_standard_qa(question, qa_pairs)
if matched:
@@ -381,18 +435,62 @@ def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any]
source = "knowledge" if references else "qwen"
config = _config(avatar)
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
chunks = _iter_qwen_stream(_build_prompt(avatar, history, question, references), temperature)
messages = _build_prompt(avatar, history, question, references)
reservation = reserve_avatar_tokens(
db,
avatar,
usage_source,
CHAT_MODEL,
messages,
CHAT_MAX_OUTPUT_TOKENS,
)
chunks = _iter_qwen_stream(messages, temperature)
if matched:
messages, reservation = [], None
if public:
source, references = "public", []
def generate():
output_parts = []
provider_usage = None
settled = False
try:
yield _sse("meta", {"source": source, "references": references})
for content in chunks:
for chunk in chunks:
if reservation is None:
content = chunk
else:
provider_usage = chunk.get("usage") or provider_usage
content = chunk.get("content")
if not content:
continue
output_parts.append(content)
yield _sse("delta", {"content": content})
yield _sse("done", {})
token_usage = None
if reservation is not None:
answer = "".join(output_parts)
token_usage = settle_reservation(
db,
reservation,
provider_usage,
fallback_total=estimate_fallback_usage(messages, answer),
)
settled = True
yield _sse("done", {} if public else {"tokenUsage": token_usage})
except RuntimeError as exc:
yield _sse("error", {"message": str(exc)})
finally:
if reservation is not None and not settled:
answer = "".join(output_parts)
if answer:
settle_reservation(
db,
reservation,
provider_usage,
fallback_total=estimate_fallback_usage(messages, answer),
)
else:
release_reservation(db, reservation, "stream_ended_without_output")
return StreamingResponse(
generate(),
@@ -441,11 +539,14 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)):
def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
avatar = _require_shared_avatar(db, share_token)
try:
result = _resolve_reply(db, avatar, body.message, body.history)
result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat")
# 公开访客无需获知知识文件名、检索分数或内部答复来源。
result["references"] = []
result["source"] = "public"
result.pop("tokenUsage", None)
return ok(result)
except InsufficientTokensError as exc:
return fail(str(exc), code=402)
except RuntimeError as exc:
return fail(str(exc), code=502)
@@ -455,15 +556,30 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N
avatar = _require_owned_avatar(db, avatar_id, authorization)
try:
return ok(_resolve_reply(db, avatar, body.message, body.history))
except InsufficientTokensError as exc:
return fail(str(exc), code=402)
except RuntimeError as exc:
return fail(str(exc), code=502)
@router.post("/avatar/{avatar_id}/chat/stream")
def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
try:
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
except InsufficientTokensError as exc:
raise HTTPException(status_code=402, detail=str(exc)) from exc
@router.post("/public/avatar/{share_token}/chat/stream")
def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
return _stream_reply(db, _require_shared_avatar(db, share_token), body.message, body.history, public=True)
try:
return _stream_reply(
db,
_require_shared_avatar(db, share_token),
body.message,
body.history,
public=True,
usage_source="public_chat_stream",
)
except InsufficientTokensError as exc:
raise HTTPException(status_code=402, detail=str(exc)) from exc
@@ -349,6 +349,9 @@ def _issue_session(db: Session, phone: str, info: dict):
db.commit()
db.refresh(user)
from services.token_billing import get_or_create_account
get_or_create_account(db, user.id)
return ok({
"token": user.app_token,
"user": user.to_dict(),
+57 -13
View File
@@ -1,37 +1,81 @@
from fastapi import APIRouter, Depends, Body
from fastapi import APIRouter, Depends, Body, Header, HTTPException
from sqlalchemy import func
from sqlalchemy.orm import Session
from database import get_db
from models import TokenAccount, TokenPlan
from models import TokenAccount, TokenPlan, TokenUsage, User
from responses import ok, fail
from services.token_billing import get_or_create_account
router = APIRouter(tags=["Token"])
def _require_user(authorization: str | None, db: Session) -> User:
if not authorization:
raise HTTPException(status_code=401, detail="未登录")
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
user = db.query(User).filter(User.app_token == token).first()
if not user:
raise HTTPException(status_code=401, detail="会话无效或已过期")
return user
@router.get("/token/balance")
def balance(db: Session = Depends(get_db)):
acc = db.query(TokenAccount).first()
return ok({"balance": acc.balance if acc else 0})
def balance(authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
acc = get_or_create_account(db, user.id)
return ok({
"balance": acc.balance,
"totalGranted": acc.total_granted,
"totalConsumed": acc.total_consumed,
})
@router.get("/token/plans")
def plans(db: Session = Depends(get_db)):
def plans(authorization: str = Header(None), db: Session = Depends(get_db)):
_require_user(authorization, db)
items = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all()
return ok([p.to_dict() for p in items])
@router.post("/token/charge")
def charge(payload: dict = Body(...), db: Session = Depends(get_db)):
def charge(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
plan_id = payload.get("planId")
plan = db.query(TokenPlan).filter(TokenPlan.id == plan_id).first()
if not plan:
return fail("套餐不存在", 404)
acc = db.query(TokenAccount).first()
if not acc:
acc = TokenAccount(balance=0)
db.add(acc)
db.commit()
db.refresh(acc)
acc = get_or_create_account(db, user.id)
acc.balance += plan.amount
acc.total_granted = int(acc.total_granted or 0) + plan.amount
db.commit()
return ok({"balance": acc.balance, "charged": plan.amount})
@router.get("/token/usage")
def usage(authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
rows = (
db.query(
TokenUsage.avatar_id,
TokenUsage.source,
func.sum(TokenUsage.prompt_tokens),
func.sum(TokenUsage.completion_tokens),
func.sum(TokenUsage.total_tokens),
func.count(TokenUsage.id),
)
.filter(TokenUsage.user_id == user.id, TokenUsage.status == "completed")
.group_by(TokenUsage.avatar_id, TokenUsage.source)
.all()
)
return ok([
{
"avatarId": avatar_id,
"source": source,
"promptTokens": int(prompt_tokens or 0),
"completionTokens": int(completion_tokens or 0),
"totalTokens": int(total_tokens or 0),
"requestCount": int(request_count or 0),
}
for avatar_id, source, prompt_tokens, completion_tokens, total_tokens, request_count in rows
])
@@ -520,7 +520,7 @@ class TakeoverService:
from routers.chat import _resolve_reply
result = _resolve_reply(db, avatar, task.prompt, history)
result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover")
answer = _plain_text_reply(result.get("answer", ""))
db.refresh(task)
if task.status != "generating":
@@ -0,0 +1,198 @@
"""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)
@@ -8,6 +8,8 @@ from models import (
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
TokenAccount,
TokenUsage,
User,
)
@@ -107,6 +109,13 @@ def authorization_context():
db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).delete(
synchronize_session=False
)
user_ids = [owner.id, other.id]
db.query(TokenUsage).filter(TokenUsage.user_id.in_(user_ids)).delete(
synchronize_session=False
)
db.query(TokenAccount).filter(TokenAccount.user_id.in_(user_ids)).delete(
synchronize_session=False
)
db.query(User).filter(User.id.in_([owner.id, other.id])).delete(
synchronize_session=False
)
@@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, patch
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from database import Base
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
@@ -59,11 +58,10 @@ class FakeBoxIM:
@pytest.fixture
def service_context():
def service_context(tmp_path):
engine = create_engine(
"sqlite://",
f"sqlite:///{tmp_path / 'takeover.db'}",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
session_factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
Base.metadata.create_all(engine)
@@ -156,7 +154,7 @@ async def test_different_contacts_generate_without_blocking_each_other(service_c
)
both_generating = Barrier(2, timeout=2)
def resolve(_db, _avatar, prompt, _history):
def resolve(_db, _avatar, prompt, _history, **_kwargs):
both_generating.wait()
return {"answer": f"回复{prompt[-1]}"}
@@ -0,0 +1,252 @@
import uuid
from concurrent.futures import ThreadPoolExecutor
from threading import Barrier
from unittest.mock import patch
import pytest
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app, seed
from models import Avatar, TokenAccount, TokenPlan, TokenUsage, User
from routers.chat import _resolve_reply, _stream_reply
from services.token_billing import (
DEFAULT_TOKEN_GRANT,
InsufficientTokensError,
get_or_create_account,
release_reservation,
reserve_avatar_tokens,
settle_reservation,
)
client = TestClient(app)
def test_balance_is_user_scoped_and_defaults_to_one_million(authorization_context):
context = authorization_context
owner = client.get("/api/token/balance", headers=context["owner_headers"])
other = client.get("/api/token/balance", headers=context["other_headers"])
assert owner.status_code == 200
assert owner.json()["data"] == {
"balance": DEFAULT_TOKEN_GRANT,
"totalGranted": DEFAULT_TOKEN_GRANT,
"totalConsumed": 0,
}
assert other.json()["data"]["balance"] == DEFAULT_TOKEN_GRANT
assert client.get("/api/token/balance").status_code == 401
def test_seed_synchronizes_requested_recharge_plans():
seed()
db = SessionLocal()
try:
plans = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all()
assert [(plan.price, plan.amount) for plan in plans] == [
(10, 2_000_000),
(100, 20_000_000),
(1000, 250_000_000),
(10000, 2_500_000_000),
]
finally:
db.close()
def test_multiple_avatars_share_owner_balance_and_usage_is_itemized():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"token-user-{suffix}", huihui_user_id=f"token-huihui-{suffix}")
first = Avatar(id=f"token-avatar-a-{suffix}", owner_id=user.huihui_user_id, name="甲")
second = Avatar(id=f"token-avatar-b-{suffix}", owner_id=user.huihui_user_id, name="乙")
db.add_all([user, first, second])
db.commit()
try:
first_reservation = reserve_avatar_tokens(db, first, "chat", "qwen-test", [{"content": "问题一"}], 128)
settle_reservation(
db,
first_reservation,
{"prompt_tokens": 60, "completion_tokens": 40, "total_tokens": 100},
fallback_total=999,
)
second_reservation = reserve_avatar_tokens(db, second, "takeover", "qwen-test", [{"content": "问题二"}], 128)
settle_reservation(
db,
second_reservation,
{"prompt_tokens": 120, "completion_tokens": 80, "total_tokens": 200},
fallback_total=999,
)
account = get_or_create_account(db, user.id)
assert account.balance == DEFAULT_TOKEN_GRANT - 300
assert account.total_consumed == 300
usages = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).order_by(TokenUsage.total_tokens).all()
assert [(row.avatar_id, row.source, row.total_tokens) for row in usages] == [
(first.id, "chat", 100),
(second.id, "takeover", 200),
]
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id.in_([first.id, second.id])).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()
def test_concurrent_settlements_do_not_overwrite_each_other():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"concurrent-user-{suffix}", huihui_user_id=f"concurrent-huihui-{suffix}")
avatar = Avatar(id=f"concurrent-avatar-{suffix}", owner_id=user.huihui_user_id, name="并发测试")
db.add_all([user, avatar])
db.commit()
first = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "甲"}], 128)
second = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "乙"}], 128)
db.close()
barrier = Barrier(2, timeout=3)
def settle(reservation, total):
thread_db = SessionLocal()
try:
barrier.wait()
settle_reservation(
thread_db,
reservation,
{"prompt_tokens": total - 20, "completion_tokens": 20, "total_tokens": total},
fallback_total=999,
)
finally:
thread_db.close()
with ThreadPoolExecutor(max_workers=2) as pool:
list(pool.map(lambda args: settle(*args), [(first, 100), (second, 200)]))
db = SessionLocal()
try:
account = get_or_create_account(db, user.id)
assert account.balance == DEFAULT_TOKEN_GRANT - 300
assert account.total_consumed == 300
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()
def test_failed_model_request_returns_the_full_reservation():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"refund-user-{suffix}", huihui_user_id=f"refund-huihui-{suffix}")
avatar = Avatar(id=f"refund-avatar-{suffix}", owner_id=user.huihui_user_id, name="退款测试")
db.add_all([user, avatar])
db.commit()
try:
reservation = reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128)
release_reservation(db, reservation, "provider error")
account = get_or_create_account(db, user.id)
usage = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).one()
assert account.balance == DEFAULT_TOKEN_GRANT
assert account.total_consumed == 0
assert usage.status == "failed"
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()
def test_insufficient_balance_rejects_before_model_usage_is_created():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"empty-user-{suffix}", huihui_user_id=f"empty-huihui-{suffix}")
avatar = Avatar(id=f"empty-avatar-{suffix}", owner_id=user.huihui_user_id, name="余额不足")
db.add_all([user, avatar])
db.commit()
try:
account = get_or_create_account(db, user.id)
account.balance = 1
db.commit()
with pytest.raises(InsufficientTokensError):
reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128)
db.refresh(account)
assert account.balance == 1
assert db.query(TokenUsage).filter(TokenUsage.user_id == user.id).count() == 0
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()
def test_chat_settles_from_provider_usage_not_fallback_estimate():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"chat-user-{suffix}", huihui_user_id=f"chat-huihui-{suffix}")
avatar = Avatar(id=f"chat-avatar-{suffix}", owner_id=user.huihui_user_id, name="聊天测试", config={})
db.add_all([user, avatar])
db.commit()
try:
with patch(
"routers.chat._call_qwen",
return_value={
"answer": "测试回答",
"usage": {"prompt_tokens": 80, "completion_tokens": 20, "total_tokens": 100},
},
):
result = _resolve_reply(
db,
avatar,
"测试问题",
[],
qa_pairs=[],
search_fn=lambda *_args: [],
)
assert result["tokenUsage"]["totalTokens"] == 100
assert result["tokenUsage"]["balance"] == DEFAULT_TOKEN_GRANT - 100
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()
@pytest.mark.asyncio
async def test_streaming_chat_settles_final_provider_usage():
suffix = uuid.uuid4().hex
db = SessionLocal()
user = User(id=f"stream-user-{suffix}", huihui_user_id=f"stream-huihui-{suffix}")
avatar = Avatar(id=f"stream-avatar-{suffix}", owner_id=user.huihui_user_id, name="流式测试", config={})
db.add_all([user, avatar])
db.commit()
try:
chunks = iter([
{"content": "流式"},
{"content": "回答"},
{"usage": {"prompt_tokens": 90, "completion_tokens": 10, "total_tokens": 100}},
])
with patch("routers.chat._iter_qwen_stream", return_value=chunks):
response = _stream_reply(db, avatar, "测试问题", [])
body = []
async for chunk in response.body_iterator:
body.append(chunk.decode() if isinstance(chunk, bytes) else chunk)
assert "流式" in "".join(body)
account = get_or_create_account(db, user.id)
usage = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).one()
assert account.balance == DEFAULT_TOKEN_GRANT - 100
assert usage.source == "chat_stream"
assert usage.total_tokens == 100
finally:
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
db.commit()
db.close()