feat(avatar): add user token accounting
This commit is contained in:
@@ -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 <> ''"
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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="")
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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()
|
||||
@@ -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<TokenBalance>('/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<TokenUsageSummary[]>('/token/usage')
|
||||
|
||||
// ==================== 授权管理 API ====================
|
||||
|
||||
export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'interact' | 'takeover'
|
||||
|
||||
@@ -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<any[]>([])
|
||||
// 全局 Token 余额(来自后端)
|
||||
// 当前用户所有分身共享的 Token 账户
|
||||
const tokenBalance = ref<number>(0)
|
||||
const tokenConsumed = ref<number>(0)
|
||||
const tokenUsageByAvatar = ref<Record<string, number>>({})
|
||||
// 当前选中分身 id
|
||||
const currentAvatarId = ref<string | null>(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<Record<string, number>>((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,
|
||||
|
||||
@@ -30,6 +30,7 @@
|
||||
<div class="token-info">
|
||||
<span class="token-label">Token 余额</span>
|
||||
<span class="token-amount">{{ tokenBalance.toLocaleString() }}</span>
|
||||
<span class="token-used">累计使用 {{ tokenConsumed.toLocaleString() }}</span>
|
||||
</div>
|
||||
<button class="recharge-btn" @click="goToRecharge">充值</button>
|
||||
</div>
|
||||
@@ -55,6 +56,7 @@
|
||||
<div class="avatar-details">
|
||||
<div class="avatar-name-row"><h2 class="avatar-name">{{ a.displayName || a.name }}</h2><span class="avatar-status"><i class="status-dot" :class="a.status"></i>{{ statusText(a.status) }}</span></div>
|
||||
<p class="avatar-desc">{{ a.description || '暂无描述' }}</p>
|
||||
<span class="avatar-token-usage">累计使用 {{ avatarTokenUsage(a.id).toLocaleString() }} Token</span>
|
||||
</div>
|
||||
</div>
|
||||
<div class="avatar-actions">
|
||||
@@ -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()
|
||||
})
|
||||
</script>
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -13,6 +13,7 @@
|
||||
<span class="balance-label">当前余额</span>
|
||||
<span class="balance-amount">{{ currentBalance.toLocaleString() }}</span>
|
||||
<span class="balance-unit">Token</span>
|
||||
<span class="balance-used">累计使用 {{ totalConsumed.toLocaleString() }} Token</span>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
@@ -28,7 +29,7 @@
|
||||
@click="selectedPlan = plan"
|
||||
>
|
||||
<div class="plan-badge" v-if="plan.badge">{{ plan.badge }}</div>
|
||||
<div class="plan-amount">{{ plan.amount.toLocaleString() }}</div>
|
||||
<div class="plan-amount">{{ formatTokenAmount(plan.amount) }}</div>
|
||||
<div class="plan-unit">Token</div>
|
||||
<div class="plan-price">¥{{ plan.price }}</div>
|
||||
<div class="plan-desc" v-if="plan.desc">{{ plan.desc }}</div>
|
||||
@@ -83,7 +84,8 @@ import { getTokenBalance, getRechargePlans, chargeToken } from '@/api'
|
||||
const router = useRouter()
|
||||
|
||||
// 当前余额
|
||||
const currentBalance = ref<number>(1250)
|
||||
const currentBalance = ref<number>(0)
|
||||
const totalConsumed = ref<number>(0)
|
||||
|
||||
// 充值套餐
|
||||
const plans = ref<Array<{
|
||||
@@ -105,6 +107,7 @@ const loadData = async () => {
|
||||
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;
|
||||
|
||||
Reference in New Issue
Block a user