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 ])