From 699bbbde57e51ad2cf293bd0671618c2f3a541b1 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 13:24:02 +0800 Subject: [PATCH] feat(avatar): add user token accounting --- digital-avatar-app/backend/database.py | 14 + digital-avatar-app/backend/main.py | 36 ++- digital-avatar-app/backend/models.py | 33 ++- digital-avatar-app/backend/routers/chat.py | 146 ++++++++-- .../backend/routers/huihui_auth.py | 3 + digital-avatar-app/backend/routers/tokens.py | 70 ++++- .../backend/services/takeover_service.py | 2 +- .../backend/services/token_billing.py | 198 ++++++++++++++ digital-avatar-app/backend/tests/conftest.py | 9 + .../backend/tests/test_takeover_service.py | 8 +- .../backend/tests/test_token_billing.py | 252 ++++++++++++++++++ digital-avatar-app/src/api/index.ts | 21 +- digital-avatar-app/src/store/avatar.ts | 22 +- digital-avatar-app/src/views/AvatarManage.vue | 18 ++ digital-avatar-app/src/views/TokenCharge.vue | 18 +- 15 files changed, 798 insertions(+), 52 deletions(-) create mode 100644 digital-avatar-app/backend/services/token_billing.py create mode 100644 digital-avatar-app/backend/tests/test_token_billing.py diff --git a/digital-avatar-app/backend/database.py b/digital-avatar-app/backend/database.py index 076e2f0..84e7d10 100644 --- a/digital-avatar-app/backend/database.py +++ b/digital-avatar-app/backend/database.py @@ -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 <> ''" + ) diff --git a/digital-avatar-app/backend/main.py b/digital-avatar-app/backend/main.py index dea64a5..57beb2a 100644 --- a/digital-avatar-app/backend/main.py +++ b/digital-avatar-app/backend/main.py @@ -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() diff --git a/digital-avatar-app/backend/models.py b/digital-avatar-app/backend/models.py index 38b1484..b4ae291 100644 --- a/digital-avatar-app/backend/models.py +++ b/digital-avatar-app/backend/models.py @@ -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="") diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index 74233e2..e01506d 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -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 diff --git a/digital-avatar-app/backend/routers/huihui_auth.py b/digital-avatar-app/backend/routers/huihui_auth.py index 25bf880..5c2ba96 100644 --- a/digital-avatar-app/backend/routers/huihui_auth.py +++ b/digital-avatar-app/backend/routers/huihui_auth.py @@ -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(), diff --git a/digital-avatar-app/backend/routers/tokens.py b/digital-avatar-app/backend/routers/tokens.py index 18871e8..3a447e9 100644 --- a/digital-avatar-app/backend/routers/tokens.py +++ b/digital-avatar-app/backend/routers/tokens.py @@ -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 + ]) diff --git a/digital-avatar-app/backend/services/takeover_service.py b/digital-avatar-app/backend/services/takeover_service.py index 43e8ad0..9a96f9d 100644 --- a/digital-avatar-app/backend/services/takeover_service.py +++ b/digital-avatar-app/backend/services/takeover_service.py @@ -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": diff --git a/digital-avatar-app/backend/services/token_billing.py b/digital-avatar-app/backend/services/token_billing.py new file mode 100644 index 0000000..ab763a1 --- /dev/null +++ b/digital-avatar-app/backend/services/token_billing.py @@ -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) diff --git a/digital-avatar-app/backend/tests/conftest.py b/digital-avatar-app/backend/tests/conftest.py index 9d2068f..53d78a8 100644 --- a/digital-avatar-app/backend/tests/conftest.py +++ b/digital-avatar-app/backend/tests/conftest.py @@ -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 ) diff --git a/digital-avatar-app/backend/tests/test_takeover_service.py b/digital-avatar-app/backend/tests/test_takeover_service.py index 5a6eed6..f8e850f 100644 --- a/digital-avatar-app/backend/tests/test_takeover_service.py +++ b/digital-avatar-app/backend/tests/test_takeover_service.py @@ -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]}"} diff --git a/digital-avatar-app/backend/tests/test_token_billing.py b/digital-avatar-app/backend/tests/test_token_billing.py new file mode 100644 index 0000000..3860a9d --- /dev/null +++ b/digital-avatar-app/backend/tests/test_token_billing.py @@ -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() diff --git a/digital-avatar-app/src/api/index.ts b/digital-avatar-app/src/api/index.ts index 59f0166..55daf23 100644 --- a/digital-avatar-app/src/api/index.ts +++ b/digital-avatar-app/src/api/index.ts @@ -131,9 +131,24 @@ export const deleteAvatar = (id: string) => // ==================== Token 管理 API ==================== +export interface TokenBalance { + balance: number + totalGranted: number + totalConsumed: number +} + +export interface TokenUsageSummary { + avatarId: string + source: string + promptTokens: number + completionTokens: number + totalTokens: number + requestCount: number +} + // 获取 Token 余额 export const getTokenBalance = () => - request.get<{ balance: number }>('/token/balance') + request.get('/token/balance') // 获取充值套餐 export const getRechargePlans = () => @@ -143,6 +158,10 @@ export const getRechargePlans = () => export const chargeToken = (planId: string) => request.post<{ balance: number; charged: number }>('/token/charge', { planId }) +// 按分身和使用场景汇总 Token 消耗 +export const getTokenUsage = () => + request.get('/token/usage') + // ==================== 授权管理 API ==================== export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'interact' | 'takeover' diff --git a/digital-avatar-app/src/store/avatar.ts b/digital-avatar-app/src/store/avatar.ts index 0af7f4f..ed303a4 100644 --- a/digital-avatar-app/src/store/avatar.ts +++ b/digital-avatar-app/src/store/avatar.ts @@ -1,13 +1,15 @@ import { defineStore } from 'pinia' import { ref } from 'vue' -import { getAvatarList, createAvatar as apiCreate, deleteAvatar as apiDelete, getTokenBalance, getUserProfile } from '@/api' +import { getAvatarList, createAvatar as apiCreate, deleteAvatar as apiDelete, getTokenBalance, getTokenUsage, getUserProfile } from '@/api' import { unwrapListData } from '@/utils/avatar-page-data' export const useAvatarStore = defineStore('avatar', () => { // 已创建的分身列表(来自后端) const avatars = ref([]) - // 全局 Token 余额(来自后端) + // 当前用户所有分身共享的 Token 账户 const tokenBalance = ref(0) + const tokenConsumed = ref(0) + const tokenUsageByAvatar = ref>({}) // 当前选中分身 id const currentAvatarId = ref(null) // 会会用户资料(头像/昵称,来自会会接口) @@ -29,11 +31,24 @@ export const useAvatarStore = defineStore('avatar', () => { try { const res = await getTokenBalance() tokenBalance.value = (res as any)?.balance ?? 0 + tokenConsumed.value = (res as any)?.totalConsumed ?? 0 } catch (e) { console.error('加载余额失败', e) } } + const loadTokenUsage = async () => { + try { + const rows = await getTokenUsage() + tokenUsageByAvatar.value = rows.reduce>((result, row) => { + result[row.avatarId] = (result[row.avatarId] || 0) + row.totalTokens + return result + }, {}) + } catch (e) { + console.error('加载 Token 用量失败', e) + } + } + // 拉取会会用户资料(头像/昵称) const loadUserProfile = async () => { // 若已通过 uniapp 壳注入(混合架构),优先保留,不回退到后端 mock @@ -82,10 +97,13 @@ export const useAvatarStore = defineStore('avatar', () => { return { avatars, tokenBalance, + tokenConsumed, + tokenUsageByAvatar, currentAvatarId, userProfile, loadAvatars, loadTokenBalance, + loadTokenUsage, loadUserProfile, setNativeProfile, addAvatar, diff --git a/digital-avatar-app/src/views/AvatarManage.vue b/digital-avatar-app/src/views/AvatarManage.vue index 117e3a2..28a5805 100644 --- a/digital-avatar-app/src/views/AvatarManage.vue +++ b/digital-avatar-app/src/views/AvatarManage.vue @@ -30,6 +30,7 @@
Token 余额 {{ tokenBalance.toLocaleString() }} + 累计使用 {{ tokenConsumed.toLocaleString() }}
@@ -55,6 +56,7 @@

{{ a.displayName || a.name }}

{{ statusText(a.status) }}

{{ a.description || '暂无描述' }}

+ 累计使用 {{ avatarTokenUsage(a.id).toLocaleString() }} Token
@@ -95,7 +97,9 @@ const me = computed(() => userStore.user) // 状态(来自 store / 后端) const tokenBalance = computed(() => avatarStore.tokenBalance) +const tokenConsumed = computed(() => avatarStore.tokenConsumed) const avatars = computed(() => avatarStore.avatars) +const avatarTokenUsage = (id: string) => avatarStore.tokenUsageByAvatar[id] || 0 const shareToast = ref('') @@ -174,6 +178,7 @@ onMounted(() => { userStore.loadFromStorage() avatarStore.loadAvatars() avatarStore.loadTokenBalance() + avatarStore.loadTokenUsage() }) @@ -316,6 +321,12 @@ onMounted(() => { color: #F97316; } +.token-used { + margin-top: 3px; + color: #A0A5B4; + font-size: 11px; +} + .recharge-btn { padding: 8px 16px; background: #F97316; @@ -440,6 +451,13 @@ onMounted(() => { white-space: nowrap; } +.avatar-token-usage { + display: inline-block; + margin-top: 5px; + color: #A0A5B4; + font-size: 10px; +} + .avatar-status { display: inline-flex; align-items: center; diff --git a/digital-avatar-app/src/views/TokenCharge.vue b/digital-avatar-app/src/views/TokenCharge.vue index f621b5c..5447e65 100644 --- a/digital-avatar-app/src/views/TokenCharge.vue +++ b/digital-avatar-app/src/views/TokenCharge.vue @@ -13,6 +13,7 @@ 当前余额 {{ currentBalance.toLocaleString() }} Token + 累计使用 {{ totalConsumed.toLocaleString() }} Token
@@ -28,7 +29,7 @@ @click="selectedPlan = plan" >
{{ plan.badge }}
-
{{ plan.amount.toLocaleString() }}
+
{{ formatTokenAmount(plan.amount) }}
Token
¥{{ plan.price }}
{{ plan.desc }}
@@ -83,7 +84,8 @@ import { getTokenBalance, getRechargePlans, chargeToken } from '@/api' const router = useRouter() // 当前余额 -const currentBalance = ref(1250) +const currentBalance = ref(0) +const totalConsumed = ref(0) // 充值套餐 const plans = ref { try { const b: any = await getTokenBalance() currentBalance.value = b?.balance ?? 0 + totalConsumed.value = b?.totalConsumed ?? 0 } catch (e) { console.error('加载余额失败', e) } @@ -118,6 +121,10 @@ const loadData = async () => { // 执行充值(写入后端) const charging = ref(false) +const formatTokenAmount = (amount: number) => { + if (amount >= 1_000_000 && amount % 1_000_000 === 0) return `${amount / 1_000_000}M` + return amount.toLocaleString() +} const doCharge = async () => { if (!selectedPlan.value || charging.value) return charging.value = true @@ -187,6 +194,7 @@ onMounted(() => { .balance-card { display: flex; + flex-wrap: wrap; align-items: baseline; gap: 8px; padding: 20px; @@ -211,6 +219,12 @@ onMounted(() => { opacity: 0.9; } +.balance-used { + flex-basis: 100%; + font-size: 12px; + opacity: 0.82; +} + /* 充值套餐 */ .plans-section { padding: 0 20px 20px;