145 lines
4.6 KiB
Python
145 lines
4.6 KiB
Python
import uuid
|
|
from datetime import datetime, timedelta
|
|
from jose import jwt, JWTError
|
|
from passlib.context import CryptContext
|
|
from fastapi import APIRouter, Depends, HTTPException, Header
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from pydantic import BaseModel
|
|
from app.database import get_db
|
|
from app.models.user import User
|
|
from app.models.game import GameSession
|
|
from app.config import settings
|
|
|
|
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
|
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
|
|
|
|
|
class RegisterRequest(BaseModel):
|
|
email: str
|
|
password: str
|
|
nickname: str
|
|
|
|
|
|
class LoginRequest(BaseModel):
|
|
email: str
|
|
password: str
|
|
|
|
|
|
class GuestRequest(BaseModel):
|
|
nickname: str = "游客"
|
|
|
|
|
|
class AuthResponse(BaseModel):
|
|
token: str
|
|
user_id: str
|
|
nickname: str
|
|
is_guest: bool
|
|
|
|
|
|
def create_token(user_id: str) -> str:
|
|
expire = datetime.utcnow() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
|
|
return jwt.encode({"sub": user_id, "exp": expire, "v": 0}, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
|
|
|
|
|
|
async def get_current_user(authorization: str = Header(""), db: AsyncSession = Depends(get_db)) -> User:
|
|
if not authorization.startswith("Bearer "):
|
|
raise HTTPException(401, "未登录")
|
|
token = authorization[7:]
|
|
try:
|
|
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
|
user_id = payload.get("sub")
|
|
except JWTError:
|
|
raise HTTPException(401, "无效的token")
|
|
result = await db.execute(select(User).where(User.id == user_id))
|
|
user = result.scalar_one_or_none()
|
|
if not user:
|
|
raise HTTPException(401, "用户不存在")
|
|
return user
|
|
|
|
|
|
@router.post("/register")
|
|
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
|
|
result = await db.execute(select(User).where(User.email == req.email))
|
|
if result.scalar_one_or_none():
|
|
raise HTTPException(400, "邮箱已注册")
|
|
user = User(
|
|
id=str(uuid.uuid4()),
|
|
email=req.email,
|
|
nickname=req.nickname,
|
|
password_hash=pwd_context.hash(req.password),
|
|
)
|
|
db.add(user)
|
|
await db.commit()
|
|
return AuthResponse(token=create_token(user.id), user_id=user.id, nickname=user.nickname, is_guest=False)
|
|
|
|
|
|
@router.post("/login")
|
|
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
|
|
result = await db.execute(select(User).where(User.email == req.email))
|
|
user = result.scalar_one_or_none()
|
|
if not user or not user.password_hash or not pwd_context.verify(req.password, user.password_hash):
|
|
raise HTTPException(401, "邮箱或密码错误")
|
|
return AuthResponse(token=create_token(user.id), user_id=user.id, nickname=user.nickname, is_guest=False)
|
|
|
|
|
|
@router.post("/guest")
|
|
async def guest_login(req: GuestRequest, db: AsyncSession = Depends(get_db)):
|
|
user = User(
|
|
id=str(uuid.uuid4()),
|
|
nickname=req.nickname,
|
|
is_guest=True,
|
|
)
|
|
db.add(user)
|
|
await db.commit()
|
|
return AuthResponse(token=create_token(user.id), user_id=user.id, nickname=user.nickname, is_guest=True)
|
|
|
|
|
|
@router.post("/logout")
|
|
async def logout():
|
|
return {"status": "ok"}
|
|
|
|
|
|
@router.get("/profile")
|
|
async def get_profile(user: User = Depends(get_current_user)):
|
|
return {
|
|
"user_id": user.id,
|
|
"email": user.email,
|
|
"nickname": user.nickname,
|
|
"avatar_url": user.avatar_url,
|
|
"is_guest": user.is_guest,
|
|
"role": user.role,
|
|
"game_count": user.game_count,
|
|
"created_at": user.created_at.isoformat() if user.created_at else "",
|
|
}
|
|
|
|
|
|
@router.put("/profile")
|
|
async def update_profile(data: dict, user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
|
|
if "nickname" in data:
|
|
user.nickname = data["nickname"]
|
|
if "avatar_url" in data:
|
|
user.avatar_url = data["avatar_url"]
|
|
await db.commit()
|
|
return {"status": "ok"}
|
|
|
|
|
|
@router.get("/history")
|
|
async def get_history(user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
|
|
result = await db.execute(
|
|
select(GameSession).order_by(GameSession.created_at.desc()).limit(50)
|
|
)
|
|
games = result.scalars().all()
|
|
return [
|
|
{
|
|
"id": g.id,
|
|
"script_id": g.script_id,
|
|
"status": g.status,
|
|
"phase": g.phase,
|
|
"started_at": g.started_at.isoformat() if g.started_at else "",
|
|
"completed_at": g.completed_at.isoformat() if g.completed_at else "",
|
|
"created_at": g.created_at.isoformat(),
|
|
}
|
|
for g in games
|
|
]
|