82 lines
2.9 KiB
Python
82 lines
2.9 KiB
Python
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, 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(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(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(...), 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 = 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
|
|
])
|